02 - 模型结构从零实现
这一篇产出一个 model.py。跑完验收脚本,它的参数量要正好是 502,193,664,前向能出 logits,初始 loss 落在 10.37 附近。
不用 from_pretrained,不用 transformers 的任何模型类。整个文件只依赖 torch 和 torch.nn,四百行以内。
需要的前置:知道矩阵乘法,写过一点 PyTorch(会 nn.Linear 和 forward 就够)。01 篇的 train.bin 这一篇用不上,纯写模型。
零、开始之前:Transformer 到底在算什么
0.1 任务还是 01 篇那个任务
再确认一遍目标,因为整个模型结构都是围着它设计的:看着前面的 token,猜下一个 token。
输入是一串 token id,输出是每个位置上「下一个 token 是词表里哪个」的概率分布。就这样。
0.2 数据在模型里走一遍
先不管内部细节,看形状怎么变。设 batch 大小为 B、序列长度为 T:
| 步骤 | 张量形状 | 在干什么 |
|---|---|---|
| 输入 | (B, T) | 整数,每个数是一个 token id |
| Embedding 查表 | (B, T, 1536) | 每个 id 换成一个 1536 维向量 |
| Block × 18 | (B, T, 1536) | 形状不变,内容被反复加工 18 次 |
| 最后的 RMSNorm | (B, T, 1536) | 归一化 |
| lm_head | (B, T, 32000) | 投影到词表大小,得到 logits |
关键在于中间 18 层形状完全不变。每一层的输入输出都是 (B, T, 1536),所以可以随便堆几层。这是 Transformer 能做深的结构性原因。
最后那个 (B, T, 32000) 叫 logits,每个位置一个长度 32000 的向量,softmax 之后就是概率分布。
0.3 一个 Block 里有什么
每层 Block 干两件事,顺序固定:
分工可以这么理解。Attention 负责 token 之间的信息交换:第 5 个位置想知道第 2 个位置说了什么,靠它。MLP 负责每个 token 自己的加工:拿到信息之后做非线性变换,它对每个位置独立操作,位置之间不通信。
那两个圆圈是残差连接,也就是 x = x + f(x)。它的作用是给梯度留一条直通的路。没有它,18 层的梯度传到第一层基本就没了。这个结构 02 篇不展开推导,记住「残差是恒等通路」就够用。
0.4 这一篇要写的五个部件
| 部件 | 作用 | 在哪一节 |
|---|---|---|
| RMSNorm | 归一化,稳住数值 | 第二节 |
| RoPE | 告诉模型 token 的位置 | 第三节 |
| GQA Attention | token 之间交换信息 | 第四节 |
| SwiGLU | 每个 token 自己加工 | 第五节 |
| Block / LLM | 把上面四个拼起来 | 第六节 |
每一节的套路都一样:先说这个部件解决什么问题、不要它会怎样,再讲它怎么做,最后给代码。
一、为什么照抄 Llama,改了 GPT-2 的哪五处
不做架构创新,结构照抄 Llama。理由很实在:出了问题可以直接跟现成实现对照排查,而且将来想加载别人的权重也方便。
相对于最经典的 GPT-2,Llama 改了五处。这五处正好就是第二到第六节的内容:
| 位置 | GPT-2 | Llama(我们用的) | 换掉的理由 |
|---|---|---|---|
| 归一化 | LayerNorm | RMSNorm | 少算一个均值,快 7% 左右,效果不掉 |
| 位置编码 | 可学习的绝对位置 | RoPE | 能表达相对位置,且能外推到更长序列 |
| 注意力 | MHA | GQA | KV cache 直接小 3 倍 |
| FFN | GELU,中间维 4d | SwiGLU,中间维 8/3 d | 同参数量下效果更好 |
| norm 位置 | Post-norm | Pre-norm | 深层训练稳定得多 |
另外所有线性层都不带 bias。原因在 2.3 节顺带说。
二、RMSNorm
2.1 为什么需要归一化
神经网络堆深了会有个麻烦:每一层的输出分布会漂。第一层输出的数值范围可能是 ±1,传到第十层可能变成 ±100,再往后可能溢出,或者反过来缩到接近 0 梯度消失。
归一化就是在每层入口把数值拉回一个稳定的范围,让后面的层总是面对差不多尺度的输入。
2.2 LayerNorm 在做什么
LayerNorm 对每个 token 的 1536 维向量,做四件事:
减均值 、除标准差 、乘一个可学习的缩放 、加一个可学习的偏置 。
注意这里是对每个 token 自己的 1536 维算均值和方差,不是跨 batch 也不是跨序列。所以它跟 batch size 无关,这点比 BatchNorm 好用。
2.3 RMSNorm 砍掉了什么
RMSNorm 的做法是把「减均值」和「加偏置」都去掉,只留缩放:
分母那一坨就是均方根(root mean square),名字由此而来。
为什么能砍掉减均值这一步,是这一节的关键。RMSNorm 原论文的观察是:LayerNorm 真正起作用的是「把向量缩放到统一尺度」这件事,而不是「把中心移到 0」。做了消融实验,去掉中心化之后效果基本不掉,但省掉了一遍求均值和一遍减法。
砍掉 bias 也是同样的道理,实测加不加差别很小,而少一组参数就少一份显存和一次加法。这也是为什么整个模型的所有 nn.Linear 都写 bias=False。
省下来的计算量不算大,但归一化在每层要做两次、18 层就是 36 次,累积起来推理能快百分之几。免费的收益没理由不要。
2.4 代码
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim)) # 就是公式里的 gamma
def forward(self, x: torch.Tensor) -> torch.Tensor:
# 统计量始终用 fp32 算,否则 BF16 下 x^2 容易损失精度
dtype = x.dtype
x = x.float()
x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
return x.to(dtype) * self.weight
有个容易忽略的细节:中间统计量要转成 fp32 再算。BF16 只有 8 位尾数,x.pow(2) 之后动态范围会被压缩得很厉害,直接在 BF16 上求和容易丢精度。转 fp32 算完再转回来,代价可以忽略,但能避免一类很难查的数值问题。
weight 初始化成全 1,也就是一开始不做任何缩放,让模型自己学。
2.5 Pre-norm 还是 Post-norm
同样是 RMSNorm,放的位置不同,训练难度差很远。
Post-norm(原始 Transformer 论文的做法)是先过子层再归一化:
x = Norm(x + Attention(x))
Pre-norm(现在几乎所有大模型的做法)是先归一化再过子层:
x = x + Attention(Norm(x))
差别在残差通路。Pre-norm 的写法里,x 是一路直接加过去的,从最后一层到第一层存在一条完全没有归一化操作的恒等通路,梯度可以无损地流回去。Post-norm 里每一层的残差都要再过一次 Norm,梯度传递会被反复缩放,层数一深就容易出问题。
代价是 Pre-norm 每层输出的方差会随层数累加,所以最后要额外补一个 RMSNorm(就是 0.2 节表里倒数第二行那个),把进 lm_head 之前的数值拉回来。
三、RoPE 旋转位置编码
3.1 Attention 根本看不见位置
这是一个不那么直观但很重要的事实:attention 本身对输入顺序是无感的。
原因在于 attention 的计算方式。 每个位置的输出是所有位置的 value 的加权和,权重由 query 和 key 的内积决定。这里面没有任何一项跟「第几个位置」有关。把输入序列打乱顺序,输出也只是跟着打乱,内容完全一样。
所以「我打你」和「你打我」在纯 attention 眼里是同一个东西。必须额外把位置信息喂进去。
3.2 早期做法和它的问题
GPT-2 的做法是搞一个可学习的位置向量表,大小是 (max_seq_len, d_model),第 t 个位置就查第 t 行,加到 embedding 上。
两个问题。一是只能处理训练时见过的长度,表只有 1024 行,第 1025 个位置查不到,模型直接不能用。二是它编码的是绝对位置,但语言里真正重要的往往是相对距离,「形容词修饰它后面第一个名词」这种规律跟这个词在全文第几个位置没关系。
3.3 RoPE 的想法:把位置变成旋转角度
RoPE 的做法很巧。它不给 embedding 加东西,而是在 attention 内部,把 query 和 key 按位置旋转一个角度。
具体做法:把 128 维的 head 向量两两分成 64 组,每组当成一个二维平面上的向量。位置 处的向量,第 组按角度 旋转:
位置越靠后,转得越多。
3.4 为什么旋转能表达相对位置
这是 RoPE 的精髓,值得推一下。
二维旋转矩阵有个性质:。也就是说一个旋转的转置乘另一个旋转,等于「差角」的旋转。
现在看 attention 里的内积。位置 的 query 和位置 的 key,各自旋转之后再做内积:
结果里只剩下 。绝对位置 和 都消掉了,只留下它们的差。
这正是我们想要的:模型感知到的是「这两个 token 隔多远」,而不是「它们分别在第几个位置」。
原理讲完要验。code/rope_demo.py 是个纯 Python 的验证,固定两个向量,固定相对距离 3,挪动绝对位置:
固定相对距离 n-m=3,挪动绝对位置:
m n n-m 旋转后内积
0 3 3 -1.7349277069
1 4 3 -1.7349277069
2 5 3 -1.7349277069
3 6 3 -1.7349277069
4 7 3 -1.7349277069
5 8 3 -1.7349277069
改变相对距离,内积才跟着变:
n-m 旋转后内积
0 -0.9800000000
1 -1.7064616876
2 -2.0008180700
3 -1.7349277069
4 -1.2924048008
5 -1.2145554743
小数点后 10 位完全一致。相对距离一变,内积立刻跟着变。
3.5 频率怎么定
64 组用同一个角度是不行的,那样只能表达一种尺度的距离。RoPE 给每组配不同的频率:
第 0 组频率是 1,转得最快,每挪一个位置就转 1 弧度,用来分辨近距离。最后一组频率是 ,转得极慢,几千个位置才转完一圈,用来分辨远距离。
这样 64 组合起来,就能同时表达从「隔 1 个」到「隔几千个」的各种距离。跟时钟的秒针分针时针是一个道理。
base 这个数还有个用处:把它调大能让所有频率变慢,等效于把位置「压缩」,这是长上下文外推最常用的手段之一。我们这篇不展开,但要知道 rope_theta 这个超参是干这个的。
3.6 代码
def build_rope_cache(seq_len, head_dim, theta, device, dtype):
"""预计算 cos / sin 表,形状都是 (seq_len, head_dim)"""
idx = torch.arange(0, head_dim, 2, device=device).float() / head_dim
inv_freq = 1.0 / (theta ** idx) # (head_dim/2,)
pos = torch.arange(seq_len, device=device).float() # (seq_len,)
freqs = torch.outer(pos, inv_freq) # (seq_len, head_dim/2)
emb = torch.cat((freqs, freqs), dim=-1) # (seq_len, head_dim)
return emb.cos().to(dtype), emb.sin().to(dtype)
def rotate_half(x):
half = x.shape[-1] // 2
x1, x2 = x[..., :half], x[..., half:]
return torch.cat((-x2, x1), dim=-1)
def apply_rope(q, k, cos, sin):
cos = cos[None, None, :, :] # 广播到 (B, n_heads, T, head_dim)
sin = sin[None, None, :, :]
return q * cos + rotate_half(q) * sin, k * cos + rotate_half(k) * sin
cos 和 sin 只跟位置有关,跟输入内容无关,所以建模型的时候算一次存下来就行,不用每次前向重算。代码里用 register_buffer 挂在模型上,它不是参数、不参与梯度,但会跟着 .to(device) 一起搬到 GPU。
3.7 一个必须知道的坑:两套约定
rotate_half 的写法是把 128 维切成前 64 和后 64,配对方式是「第 0 维配第 64 维」。这是 HuggingFace 的约定。
RoPE 原论文用的是另一套:相邻两维配对,也就是「第 0 维配第 1 维」。
两套都能正常训练,效果也一样,但它们互不兼容。 用 A 约定训出来的权重,拿 B 约定的代码加载,前向输出全是乱的,而且不会报任何错,只会表现为模型胡说八道。
这个坑在 08 篇会真的碰上,因为那时候要把权重喂给自制推理框架。训练端和推理端必须用同一套约定,现在选了 rotate_half,到时候推理框架也得是 rotate_half。
四、GQA 分组查询注意力
4.1 先说 MHA
标准多头注意力(MHA)的做法:把 1536 维切成 12 个头,每个头 128 维,各自独立算一遍 attention,最后拼回来。
每个头都有自己的 Q、K、V 三个投影。多个头的意义是让不同的头关注不同类型的关系,有的头管语法,有的头管指代。
4.2 推理时的真正瓶颈是 KV cache
训练的时候一次算整个序列,MHA 没什么问题。问题出在推理。
推理是一个 token 一个 token 往外吐的。生成第 100 个 token 时,需要用到前面 99 个 token 的 K 和 V。如果每次都重算,复杂度是平方级的,慢得没法用。
所以实际做法是把算过的 K 和 V 缓存起来,这就是 KV cache。每生成一个 token,就往缓存里追加一份 K 和 V。
问题是这个缓存很占显存,而且跟并发数成正比。服务 100 个用户就要 100 份 KV cache。在推理服务里,KV cache 经常比模型权重本身还占地方。vLLM 的 PagedAttention 解决的就是这块的碎片问题(见 PagedAttention)。
4.3 MQA 和 GQA
既然 KV cache 是瓶颈,一个自然的想法是:能不能让多个头共享同一份 K 和 V?
MQA(Multi-Query Attention)走到极端,所有 Q 头共享唯一一组 K/V。缓存直接小 12 倍,但效果掉得比较明显。
GQA(Grouped-Query Attention)取中间:把 Q 头分组,每组共享一组 K/V。我们的配置是 12 个 Q 头、4 个 KV 头,也就是每 3 个 Q 头共享 1 组 K/V。
4.4 我们这个配置省了多少
每个 token 要缓存的字节数:
开头的 2 是因为 K 和 V 各存一份。代进去(BF16 所以每个数 2 字节):
| 方案 | n_kv_heads | 每 token | 2048 长的一条序列 |
|---|---|---|---|
| MHA | 12 | 108.0 KB | 226.5 MB |
| GQA(我们的) | 4 | 36.0 KB | 75.5 MB |
省了整整 3 倍。 这个数字是 code/count_params.py 算出来的。
同时注意 4.1 节说的:Q 头还是 12 个,所以模型的表达能力基本没损失,掉的只是 K/V 的多样性。这是 GQA 的性价比所在,现在从 Llama 2 70B 到 Qwen 系列基本都用它。
顺带看参数量。因为 K/V 投影变窄了,Attention 的参数也跟着少了:
| 投影 | 形状 | 参数量 |
|---|---|---|
q_proj | 1536 → 12×128 = 1536 | 2,359,296 |
k_proj | 1536 → 4×128 = 512 | 786,432 |
v_proj | 1536 → 4×128 = 512 | 786,432 |
o_proj | 1536 → 1536 | 2,359,296 |
| 合计 | 6,291,456 |
如果是 MHA,k/v 各是 2,359,296,Attention 合计会变成 9,437,184。
4.5 代码
Q 头和 KV 头数量对不上,算 attention 之前要把 KV 复制对齐:
def repeat_kv(x, n_rep):
"""(B, n_kv_heads, T, hd) -> (B, n_kv_heads * n_rep, T, hd)"""
b, n_kv, t, hd = x.shape
if n_rep == 1:
return x
return (x[:, :, None, :, :]
.expand(b, n_kv, n_rep, t, hd)
.reshape(b, n_kv * n_rep, t, hd))
这里用 expand 而不是 repeat,因为 expand 不真的复制数据,只是改 stride。真正的 KV cache 里只存 4 份,复制只发生在计算的那一刻。 要是这里写成 repeat,显存收益就没了。
完整的 Attention:
class Attention(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.n_heads, self.n_kv_heads = c.n_heads, c.n_kv_heads
self.head_dim = c.head_dim
self.n_rep = c.n_heads // c.n_kv_heads # 12 // 4 = 3
self.q_proj = nn.Linear(c.d_model, c.n_heads * c.head_dim, bias=False)
self.k_proj = nn.Linear(c.d_model, c.n_kv_heads * c.head_dim, bias=False)
self.v_proj = nn.Linear(c.d_model, c.n_kv_heads * c.head_dim, bias=False)
self.o_proj = nn.Linear(c.n_heads * c.head_dim, c.d_model, bias=False)
def forward(self, x, cos, sin):
B, T, _ = x.shape
q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(B, T, self.n_kv_heads, self.head_dim).transpose(1, 2)
q, k = apply_rope(q, k, cos[:T], sin[:T]) # RoPE 只作用于 q 和 k
k, v = repeat_kv(k, self.n_rep), repeat_kv(v, self.n_rep)
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
out = out.transpose(1, 2).contiguous().view(B, T, -1)
return self.o_proj(out)
两个要点。
RoPE 只加在 q 和 k 上,不加在 v 上。 因为位置信息是通过内积起作用的(3.4 节那个推导),而 v 不参与内积,它只是被加权求和。给 v 加旋转没有意义。
is_causal=True 自带因果掩码。 因果的意思是第 t 个位置只能看到 0 到 t,不能看到未来。F.scaled_dot_product_attention 会自动生成下三角掩码,而且它会走 FlashAttention 的融合实现,比自己建一个 (T, T) 的 mask 矩阵又快又省显存。序列长 2048 时,显式 mask 光自己就要 8 MB。
4.6 顺带说 head_dim 为什么是 128
head_dim = d_model / n_heads = 1536 / 12 = 128。
128 是个几乎所有主流模型都在用的数。原因是 FlashAttention 这类融合 kernel 是按 64 / 128 这些尺寸做了专门优化的,取别的值会掉到通用实现上,慢不少。定超参的时候,让 head_dim 落在 64 或 128 是条实用的约束。
五、SwiGLU
5.1 标准 FFN 长什么样
原始 Transformer 的前馈层就两层线性加一个激活:
def ffn(x):
return W2(gelu(W1(x))) # 1536 -> 4096 -> 1536
先升维再降维,中间那层一般取 4 倍宽。它对每个位置独立操作,是模型存储「知识」的主要地方。
5.2 门控的想法
SwiGLU 的改动是加一条门控支路:
是逐元素相乘。直观理解:up 那一路算出候选值,gate 那一路算出「每一维该放行多少」,两者相乘。SiLU 的输出可以接近 0,等于把某些维度关掉。
相比固定的激活函数,门控让「哪些信息通过」变成数据相关的,模型能学得更细。
这类改动 Noam Shazeer 那篇论文(GLU Variants Improve Transformer)做了一组消融,结论是同参数量下 SwiGLU 稳定地好一点。没有特别深刻的理论解释,论文里那句结论大意是把它归功于运气。工程上照用就是了。
5.3 为什么中间维是 8/3 d,不是 4 d
这是个容易被忽略但很实际的问题。
标准 FFN 两个矩阵,参数量是 。取 时是 。
SwiGLU 有三个矩阵,参数量是 。要让参数量保持在 不变:
所以 ,正好是整数,也正好是 2 的幂,对齐得很舒服。多一个矩阵就把中间维按比例调窄,总参数量不变,这样跟标准 FFN 比才是公平的。
5.4 代码
class SwiGLU(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.gate_proj = nn.Linear(c.d_model, c.ffn_dim, bias=False)
self.up_proj = nn.Linear(c.d_model, c.ffn_dim, bias=False)
self.down_proj = nn.Linear(c.ffn_dim, c.d_model, bias=False)
def forward(self, x):
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
六、组装
6.1 Block
按 0.3 节那张图,Pre-norm 加两条残差:
class Block(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.attn_norm = RMSNorm(c.d_model, c.norm_eps)
self.attn = Attention(c)
self.mlp_norm = RMSNorm(c.d_model, c.norm_eps)
self.mlp = SwiGLU(c)
def forward(self, x, cos, sin):
x = x + self.attn(self.attn_norm(x), cos, sin)
x = x + self.mlp(self.mlp_norm(x))
return x
注意 x = x + f(norm(x)),不是 x = f(norm(x))。少写那个 x + 就没有残差了,模型能跑但训不深。
6.2 完整模型
class LLM(nn.Module):
def __init__(self, c: Config):
super().__init__()
self.config = c
self.embed = nn.Embedding(c.vocab_size, c.d_model)
self.blocks = nn.ModuleList(Block(c) for _ in range(c.n_layers))
self.norm = RMSNorm(c.d_model, c.norm_eps)
self.lm_head = nn.Linear(c.d_model, c.vocab_size, bias=False)
if c.tie_embeddings:
self.lm_head.weight = self.embed.weight
cos, sin = build_rope_cache(c.max_seq_len, c.head_dim, c.rope_theta,
device="cpu", dtype=torch.float32)
self.register_buffer("cos", cos, persistent=False)
self.register_buffer("sin", sin, persistent=False)
def forward(self, idx, targets=None):
B, T = idx.shape
x = self.embed(idx)
cos, sin = self.cos[:T].to(x.dtype), self.sin[:T].to(x.dtype)
for blk in self.blocks:
x = blk(x, cos, sin)
x = self.norm(x)
logits = self.lm_head(x)
if targets is None:
return logits, None
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.reshape(-1))
return logits, loss
persistent=False 的意思是这两个 buffer 不存进 checkpoint。它们是纯计算出来的,加载时重新算一遍就行,没必要占 checkpoint 的体积。
6.3 权重共享
self.lm_head.weight = self.embed.weight 这一行让输入的 embedding 表和输出的投影矩阵用同一组权重。
两者形状正好都是 (32000, 1536),物理上可以共享。直觉上也说得通:embedding 是「token id 到语义向量」,lm_head 是「语义向量到 token id」,互为逆向,共享一组参数是合理的。
收益是省 49,152,000 个参数,占总量 9.8%。对小模型来说这个比例不小,所以小模型基本都开权重共享,大模型反而常常不开,因为那时候 embedding 占比已经很低了。